前兩天我們做了兩件事:
Day 6:把 Phase 22 拆成三段 pipeline
Day 7:把 CRT 表和 reconstruct kernel 講清楚
今天進到真正會乘矩陣的地方。
但先說結論:今天不追速度。
Day 8 只問一句:
7 張 INT8 餘數矩陣,最笨的乘法怎麼寫才是對的?
因為只要餘數矩陣錯,或 naive INT8 GEMM 錯,後面換成 Tensor Core 只會把錯誤算得更快。
Phase 22 的輸入是 FP64:
A: double[m, k]
B: double[k, n]
Tensor Core 不能直接吃 double,所以第一步是 Ozaki split。
程式用兩個 scale:
sh = 2^17
sl = 2^34
對每個元素 v,先取 high:
hi = round(v * sh)
再把 high 拿掉,取 low:
lo = round((v - hi / sh) * sl)
這裡的 hi、lo 是整數。最後重建 double 時才乘回:
hi 對應 2^-17
lo 對應 2^-34
所以兩個矩陣相乘會有四個 pass:
A_hi × B_hi
A_hi × B_lo
A_lo × B_hi
A_lo × B_lo
對應縮放:
2^-34, 2^-51, 2^-51, 2^-68
今天先拿其中一個 pass 想就好。假設我們要算 A_hi × B_hi。
precompute_modulo:把薄片變成 7 份餘數high/low 整數還是太大。
例如 v 在 [-1, 1] 附近時:
round(v × 2^17) 可能接近 ±131072
這超過 INT8。
所以對每個質數取餘數:
const int pr[7] = {127, 113, 109, 107, 103, 101, 97};
for (int p = 0; p < 7; ++p) {
int32_t rem = iv % pr[p];
if (rem < 0) rem += pr[p];
d[p * padded_size + out_off] = (int8_t)rem;
}
這就是 precompute_modulo_hl_A 和 precompute_modulo_hl_B 在做的事。
它們一次輸出兩份:
d_h:high slice 的餘數
d_l:low slice 的餘數
而且每份都有 7 個質數。
所以 A 會變成:
A_hi_mod[7]
A_lo_mod[7]
B 也一樣:
B_hi_mod[7]
B_lo_mod[7]
這些才是真正會被 INT8 GEMM 讀進去的矩陣。
Day 3 的 naive GEMM 可以直接把矩陣想成 row-major:
A[row, k]
B[k, col]
但 Phase 22 後面要餵 mma.sync.m16n8k32.s8,所以 precompute_modulo 不是單純照 row-major 寫出來。
A 被排成:
128 × 32 tile
B 被排成:
64 × 32 tile
原因是後面的 C tile 是:
128 × 64
K 方向每次吃:
32
也就是:
A tile: 128 rows × 32 K
B tile: 64 cols × 32 K
C tile: 128 rows × 64 cols
這個 layout 不是為了讓 Day 8 naive 版本舒服,而是為了 Day 9 的 Tensor Core 版本。
Day 8 如果只是寫一顆教學用 naive GEMM,可以先用正常 row-major 的 A8、B8 對答案。等答案對了,再把同一件事搬到 tiled layout。
工程上就是兩層:
正確性層:
A8[row, k], B8[k, col]
效能層:
A8 tiled 128×32, B8 tiled 64×32
不要一開始就把兩層混在一起。
最小版本跟 Day 3 的 naive GEMM 幾乎一樣,只是輸出先不要急著變成 double。
對 CRT 來說,單一質數的 GEMM 產物應該先是一張 int32 accumulator:
__global__ void naive_int8_residue_gemm(
const int8_t* A8,
const int8_t* B8,
int32_t* Cp,
int m, int n, int k)
{
int row = blockIdx.y * blockDim.y + threadIdx.y;
int col = blockIdx.x * blockDim.x + threadIdx.x;
if (row < m && col < n) {
int32_t sum = 0;
for (int kk = 0; kk < k; ++kk) {
sum += (int32_t)A8[row * k + kk]
* (int32_t)B8[kk * n + col];
}
Cp[row * n + col] = sum;
}
}
一個 thread 算一個 C[i,j]。
這顆 kernel 有三個特點:
慢的原因也很明顯:每個 thread 自己跑完整個 K 迴圈,沒有 tiling,沒有 shared memory reuse,也沒有 Tensor Core。
但它有一個很重要的用途:確認 A8 @ B8 這件事本身是對的。
等 7 個質數的 Cp 都算完,下一步才是:
v_p = Cp mod p
x = CRT(v_127, v_113, ..., v_97)
C += x × scale
也就是說,scale 不屬於單一質數的 naive GEMM。它屬於 CRT 重建後的 Ozaki pass。
今天最容易犯的錯,是一開始就拿最後的 double C 去比 cublasDgemm。
那樣錯了也很難查。
Day 8 應該分層對:
拿一個 double:
v = -0.123456
CPU 算:
hi = round(v × 2^17)
lo = round((v - hi / 2^17) × 2^34)
再算:
hi mod 127
hi mod 113
...
GPU 的 precompute_modulo_hl_A/B 輸出要一樣。
固定 p = 127。
CPU 算:
C_127 = A_127 @ B_127
C_127 = C_127 mod 127
GPU naive INT8 GEMM 的 sum 再 mod 127,要對得上。
把 7 個 C_p 丟給 Day 7 的 reconstruct:
x = Σ (C_p mod p) × f_p mod M
看它能不能拼回原本那個整數內積。
最後才做:
C += x × 2^-34
這時才拿 double 去跟 CPU 或 native DGEMM 比。
這樣切開以後,錯誤會很好定位:
單元素錯 → precompute_modulo
單質數 GEMM 錯 → INT8 multiply / layout
CRT 錯 → coeff / signed fold
double 錯 → scale / pass mapping
% 是 Day 8 最大的 debug 點昨天講過一次,今天會真的踩到。
C++ 的:
iv % p
如果 iv 是負數,結果也可能是負數。
但我們要塞進 int8_t 的是 [0, p)。
所以一定要:
int32_t rem = iv % p;
if (rem < 0) rem += p;
同樣地,INT8 GEMM 的 accumulator 也可能是負的嗎?
在 Phase 22 的餘數表示下,輸入餘數是 0..126,所以單次乘法是非負,sum 也是非負。這條 fast path 裡可以用 uint32_t 的 magic remainder。
但如果你用別的表示法,例如把餘數中心化成 [-p/2, p/2],那 accumulator 就會有負數。這時 reconstruct 前又要折回 [0, p)。
所以 Day 8 的建議是:先用最無聊的表示法。
所有餘數都存 [0, p)
所有 CRT 輸入都先保證 [0, p)
等答案穩了,再談有號表示能不能省一點東西。
今天這顆 naive INT8 GEMM 不走 Tensor Core。
它用的是普通 CUDA core:
load int8
轉 int32
乘法
加到 int32 sum
每個 thread 算一個輸出,K 次迴圈全部自己做。A、B 的資料也沒有被 block 內其他 thread 好好重用。
更麻煩的是,Phase 22 不是只做一次 GEMM:
4 個 Ozaki pass × 7 個質數 = 28 次 INT8 GEMM
如果每次都 naive,速度一定炸掉。
所以 Day 8 的 naive kernel 不是最終實作,它是測試儀器。
它回答的是:
我們餵給 Tensor Core 的東西,數學上到底對不對?
Day 9 才回答:
怎麼讓同一件事真的走到 INT8 Tensor Core?
現在 Phase 22 fast path 沒有真的呼叫 Day 8 這顆 naive kernel。
它做的是:
precompute_modulo_hl_A/B
↓
crt_fast_kernel
crt_fast_kernel 裡面已經包含:
讀 tiled A8/B8
cp.async 搬到 shared memory
ldmatrix 讀 fragment
mma.sync 做 INT8 GEMM
CRT reconstruct
寫回 double C
Day 8 把「INT8 GEMM」單獨拿出來,是為了把正確性講清楚。
更精確地說,Day 8 是 fast path 的參考答案。
fast path 做了很多效能相關的事情:
global memory layout 改排
shared memory staging
cp.async pipeline
ldmatrix fragment load
mma.sync accumulator
CRT inline reconstruct
任何一個環節錯了,最後的 double C 都會錯。
所以開發時不要直接拿 fast path 跟 cublasDgemm 硬比。比較好的方式是讓 Day 8 的 naive path 先產生一份中間答案,然後一層一層對:
precompute_modulo output
↓
單一 prime 的 int32 accumulator
↓
accumulator mod p
↓
7-prime CRT reconstruct
↓
乘回 Ozaki scale 後的 double C
fast path 至少要能在這些切點上對到 naive path。這比最後只看 max error 有用很多。
尤其 layout 這一段最容易出錯。
Day 8 的教學版可以先假設:
A8[row, kk]
B8[kk, col]
但 fast path 讀的是 tiled layout。這代表同一個數學元素在 memory 裡的位置可能完全不同。debug 時我會先準備一個很小、很容易辨認的矩陣:
A[row, kk] = 1000 * row + kk
B[kk, col] = 1000 * kk + col
然後只看 precompute_modulo 後的某一個 prime,例如 p = 127。如果某個位置本來應該是:
(1000 * row + kk) mod 127
但 device buffer 裡讀到別的值,就先不要碰 GEMM。這通常表示 layout mapping 或 padding offset 已經錯了。
padding 也要特別看。
因為 Tensor Core path 會把 M、N、K 補到 tile 友善的大小,實際 buffer 裡可能有一部分是 padding。padding 區域必須是 0,不能留下未初始化資料。否則 naive row-major 測試可能沒事,但 tiled fast path 在 K tail 或邊界 tile 會把垃圾值乘進 accumulator。
這類 bug 的特徵是:
小尺寸剛好整除 tile 時正常
一換成非 128/64/32 倍數就爆
錯誤集中在矩陣右邊或下邊
如果看到這種 pattern,先查 padding,不要先改 CRT。
repo 裡也有獨立的 INT8 GEMM kernel:
int8_gemm_s8s8s32(...)
它是後來把 s8×s8→s32 GEMM 孤立出來的版本,使用 128×128 或 128×256 tile、cp.async、ldmatrix、mma.sync。那已經比較接近 Day 9/10 的內容。
今天先不要跳到那裡。
先把最笨的 A8 @ B8 寫對。
case_2/docs2 和 case_2/hand_over_doc 裡記了不少實際 debug 問題。這些東西不一定適合全部塞進 Day 8 正文,但真的建置和測試時很有幫助。
我整理成一篇附錄:
Blog/debug_notes.md
跟 Day 8 最相關的是這幾個:
Phase22,不要不小心跑到 Phase24ExtremeMix。enable_fp64 和 enable_residual 要是 false,否則你可能在看 residual FP64 的結果。CUBLAS_EMULATE_DOUBLE_PRECISION 這類環境變數會污染 native / emu 對照。Day 8 的重點是「每一層都能對答案」。這些 handover 筆記就是在提醒:很多 bug 看起來像數學錯,其實是環境、layout 或同步問題。
我會把實際 debug 順序寫得更死一點。
第一步,只測 precompute_modulo_hl_A/B。
輸入不要先用 random。先用固定小矩陣,手動包含:
0
1
-1
127
-127
128
-128
一個普通小數,例如 0.125
一個負小數,例如 -0.125
這樣可以一次抓三種錯:
round / scale 錯
負數 modulo 錯
int8 cast 錯
第二步,只測單一 prime 的 GEMM。
先固定 p = 127,只看:
A_127 @ B_127
這一步不要做 CRT,也不要乘回 double scale。輸出就是 int32 accumulator,然後 CPU 也用同一批 A_127、B_127 算一份。兩邊 accumulator 對上後,再比較 acc % 127。
如果 accumulator 對不上,問題通常在:
A/B layout
K 迴圈範圍
padding 是否清 0
row/col index mapping
如果 accumulator 對得上,但 acc % p 對不上,才去查 remainder implementation。
第三步,才測 7 個 prime 的 CRT。
這一步要刻意測負的整數內積。例如選 A 和 B,讓某個 C[i,j] 是負數。因為 CRT reconstruct 會先得到 [0, M) 裡的代表值,再靠:
if x > M/2:
x = x - M
折回 signed integer。沒有測負內積,就等於沒有測這條 signed fold。
第四步,才測 Ozaki 四個 pass 的 scale。
四個 pass 的 scale 不一樣:
hi × hi: 2^-34
hi × lo: 2^-51
lo × hi: 2^-51
lo × lo: 2^-68
如果 CRT 的整數內積都對,但 double 結果還是錯,就要查 pass mapping 和 scale。這時候不要回頭亂改 prime 表。
最後才跑大尺寸 random matrix。
random matrix 很適合找殘留 bug,但不適合作為第一個測試。它的問題是資訊量太低:你只會知道 max error 很大,卻不知道錯在哪一層。Day 8 的做法是先把每一層都變成可以被單獨觀察的東西,讓後面的 Tensor Core 優化有一個可信的基準。
Day 8 結束時,應該有四個東西:
precompute_modulo 可以把 high/low 轉成 7 份 int8_t 餘數。A8 @ B8 對上。這時候它一定很慢。
慢是正常的。
今天的成功標準不是 ms,而是:
每一層都能對答案
錯了知道是哪一層
今天可以帶走三句:
precompute_modulo 是 double 到 INT8 世界的入口。 它把 high/low 整數薄片轉成 7 份餘數矩陣。C[i,j],很慢,但最好查。明天把同一個 A8 @ B8 搬上 Tensor Core:128×64 tile、K=32、cp.async、ldmatrix,最後用 mma.sync.m16n8k32.s8 真的把 INT8 乘法吃起來。